iT邦幫忙

2026 iThome 鐵人賽

DAY 16
0
Software Development

GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)系列 第 16

# Day 16:先把 Hadamard 塞進一個 block ( (重賽版))

  • 分享至 

  • xImage
  •  

昨天我們把題目換掉了(換題目水 嘿嘿)。

一樣 先ai slob 然後我再看時間修 。 總之 先勉強看個

這次不是 GEMM,而是:

y = H^{\otimes N} x

如果把 H^ N 展開成矩陣,那會是一個非常大的 dense matrix;但真正要算的不是 dense GEMM,而是一層一層的 butterfly。

今天先看第一段切分:當 N 還不大時,能不能把整個 FWT chunk 放在一個 CUDA block 裡面做完?

答案是可以,但 kernel 不能寫成「每個 stage 都讀寫 global memory」。那樣會把 FWT 最糟糕的一面放大:每一層只做一次加減,卻一直搬資料。

所以這個版本的方向是:

global memory 只進出一次
中間的 butterfly 盡量留在 register / warp / shared memory

一個 block 要吃多少資料

FWT 的 vector 長度是 2^N

假設用 FP32,一個 element 是 4 bytes:

N=10: 2^10 floats = 4 KB
N=12: 2^12 floats = 16 KB
N=14: 2^14 floats = 64 KB
N=15: 2^15 floats = 128 KB

這就是 single-block kernel 的第一個邊界。

N=14 時,整個 chunk 是 64 KB。這已經很接近 shared memory 壓力,但還能靠 register blocking 和 shared memory exchange 撐住。

N=15 時,整個 chunk 是 128 KB。如果你想把完整資料丟進 shared memory,會很不舒服,甚至直接撞限制。

所以真正的寫法不是:

把 128 KB state 全部放進 shared memory

而是:

每個 thread 拿一部分資料放 register
shared memory 只拿來做跨 warp 的資料重排

這是 Day 16 的核心。


kernel 的資料流

case_3 裡的 native path 主要在:

src/ImplicitHadamardNative.cu

single-block 的核心 kernel 是:

template<int N, int THREADS, int NUM_B>
__global__ void native_fp32_extreme_fwt_kernel(
    const float* __restrict__ in_state,
    float* __restrict__ out_state,
    float norm_factor
)

這裡的 NUM_B 很重要。

每個 thread 不只處理一個 float,而是處理多個 float4

float4 reg[NUM_B];

也就是:

每個 thread 持有 NUM_B 組 float4
每組 float4 = 4 個 FP32

所以一個 block 處理的元素數量是:

THREADS * NUM_B * 4

例如 N=14

THREADS = 256
NUM_B   = 16

256 * 16 * 4 = 16384 = 2^14

這剛好是一整個 N=14 chunk。


Stage 0 和 Stage 1:先在 thread 裡面做掉

每個 float4 裡面有 4 個值:

v.x, v.y, v.z, v.w

FWT 的前兩層 stride 很小,所以不需要跨 thread。

可以直接在同一個 thread 裡做:

float a0 = reg[b].x + reg[b].y;
float b0 = reg[b].x - reg[b].y;
float a1 = reg[b].z + reg[b].w;
float b1 = reg[b].z - reg[b].w;

reg[b].x = a0 + a1;
reg[b].y = b0 + b1;
reg[b].z = a0 - a1;
reg[b].w = b0 - b1;

這裡沒有 shared memory,也沒有 warp shuffle。

資料在 register 裡,ALU 直接做加減。

這也是為什麼一開始用 float4 很自然:它不是為了好看,而是剛好可以把前兩層 butterfly 吃掉。


Stage 2 到 Stage 6:warp shuffle

接著 stride 會跨 thread,但還在同一個 warp 裡。

FWT 的配對規則是:

peer = lane xor stride

所以 CUDA 裡可以直接用:

float peer = __shfl_xor_sync(0xffffffff, value, stride);

float4 來說,就是四個 lane value 都各自 shuffle:

peer.x = __shfl_xor_sync(0xffffffff, reg[b].x, stride_thread);
peer.y = __shfl_xor_sync(0xffffffff, reg[b].y, stride_thread);
peer.z = __shfl_xor_sync(0xffffffff, reg[b].z, stride_thread);
peer.w = __shfl_xor_sync(0xffffffff, reg[b].w, stride_thread);

然後根據目前 lane 在 pair 的上半部或下半部,決定是:

上半部:self + peer
下半部:peer - self

這段最關鍵的地方是:warp shuffle 是 register-to-register 的資料交換。

它沒有把資料寫回 global memory。

對 FWT 這種每次只做 add/sub 的 workload,這比多幾個 FLOPS 更重要。


跨 warp:shared memory 不是拿來算,是拿來換位置

到了更大的 stride,資料會跨 warp。

warp shuffle 只能處理同一個 warp 的 lane exchange,所以這時候需要 shared memory。

但這裡要小心:shared memory 不應該變成主要計算場所。

比較好的用法是:

register -> shared memory 做 layout exchange
shared memory -> register
接著繼續用 warp shuffle / register butterfly

這就是後來 N=15 版本最重要的洞見。

原本直覺可能會覺得:

N=15 有 128 KB,所以需要 128 KB shared memory

但實際上可以改成:

128 KB 的資料分散在每個 thread 的 register
shared memory 只拿來做 corner turning

這個差別非常大。


N=15 的版本為什麼不一樣

case_3 裡後來有一個更激進的版本:

native_fp32_extreme_fwt_kernel_v2<15, 256, 32>

這個配置代表:

THREADS = 256
NUM_B   = 32

256 * 32 * 4 = 32768 = 2^15

也就是一個 block 直接吃完整個 N=15 chunk。

每個 thread 持有:

32 個 float4 = 128 個 float

這會帶來很高的 register 壓力,但它換來一件事:

N=15 不需要切成兩個 kernel
中間不需要完整寫回 global memory 再讀回來

這正是 N=15 能不能快的關鍵。


為什麼 shared memory 要 swizzle

當資料跨 warp 交換時,如果大家用很規則的方式寫 shared memory,很容易打到同一批 bank。

shared memory bank conflict 會讓看起來很近的資料交換變慢。

所以 N=15 版本用了 XOR swizzle 的概念:

smem index 裡混入 lane_id xor warp_id

概念上就是讓不同 warp/lane 的資料不要全部撞到同一個 bank。

文章先不把 swizzle 寫成完整推導;這部分後面如果要單獨拆,可以再把 bank mapping 畫出來。

今天只需要掌握一件事:

shared memory 在這裡不是 cache,
而是一個可控的 cross-warp exchange buffer。

dispatch 怎麼選

native path 的 dispatch 大致像這樣:

N=7  -> native_fp32_extreme_fwt_kernel<7,  32,  1>
N=8  -> native_fp32_extreme_fwt_kernel<8,  64,  1>
N=9  -> native_fp32_extreme_fwt_kernel<9, 128,  1>
N=10 -> native_fp32_extreme_fwt_kernel<10,256,  1>
N=11 -> native_fp32_extreme_fwt_kernel<11,256,  2>
N=12 -> native_fp32_extreme_fwt_kernel<12,256,  4>
N=13 -> native_fp32_extreme_fwt_kernel<13,256,  8>
N=14 -> native_fp32_extreme_fwt_kernel<14,256, 16>
N=15 -> native_fp32_extreme_fwt_kernel_v2<15,256,32>

這個表其實就是一個容量表。

N 每增加 1,資料量翻倍。當 thread 數固定在 256 時,就只能讓每個 thread 拿更多 float4

所以 NUM_B 會一路增加:

1, 2, 4, 8, 16, 32

這不是 tuning magic,而是 2^N 的結果。


正確性怎麼驗

這種 kernel 最容易出現的錯不是 crash,而是 silent wrong answer。

尤其是 N=15 的 exchange,如果 swizzle 或讀回順序錯了,程式照樣跑完,但 butterfly pair 已經亂掉。

所以驗證不能只看 norm 差不多。

至少要做:

1. 小 N 對 naive PyTorch FWT
2. N=15 對 Dao-AILab fast-hadamard-transform
3. 檢查 normalization 是否一致
4. 看 max diff,而不是只看平均誤差

case_3 的結果裡,N=15 有一個很重要的數字:

QDP Native FP32: 82.98 us
Dao FHT:         98.35 us
max diff:        0.0e+00

0.0e+00 不是說 floating point 永遠沒有誤差,而是代表在這個 benchmark 設定下,計算順序跟比較標的對齊了。

這對 FWT 很重要,因為不同 butterfly 順序會讓 FP32 rounding path 不同。


Day 16 的結論

single-block FWT 的重點不是把算力打滿。

重點是讓資料不要離開晶片:

thread-local:float4 內部 butterfly
warp-local:__shfl_xor_sync
cross-warp:shared memory swizzle
large stage:register blocking

N=14 是比較自然的 shared-memory/register 邊界。

N=15 則需要更激進:每個 thread 拿 32 個 float4,用 register space 承擔完整 chunk,再用 shared memory 做 corner turning。

下一篇要看的是:如果 N > 15,一個 block 放不下了,這時候要怎麼切。


上一篇
Day 15:不是所有東西都該丟進 Tensor Core (重賽版)
下一篇
Day 17:一個 block 放不下之後,FWT 要怎麼切 [重賽版]
系列文
GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)17
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言